% Fully connected feedforward network + GWO
clc
clear all
%  LOAD DATA  
load Gammapre_Noneinput
load output_YBOCS_None
load idx_GammaNone3
input=Gammapre_Noneinput(:,[8,3,4]);
target1=output_YBOCS_None; 

 K = 12;
 X_train=input';
 Y_train=target1;
 Y_train1=Y_train';
%   Y_train=Y_train';

% CREATE THE NEURAL NETWORK  
hiddenLayerSize = 10; % Define the size of the hidden layer  
final_net = feedforwardnet(hiddenLayerSize);  

% CONFIGURE THE NETWORK WITH TRAINING DATA  
final_net = configure(final_net, X_train, Y_train);  % Transpose if needed  

% CHECK NUMBER OF WEIGHTS AND BIASES  
totalWeights = numel(getwb(final_net));  
disp(['Total Weights after configuration: ', num2str(totalWeights)]);  % This should not be 0  

% RUN GWO TO FIND THE BEST WEIGHTS  
num_iterations = 1500; % Define the number of iterations  
num_wolves = 50; % Define the number of wolves in the GWO  
alpha_pos = gwo_optimization(final_net, X_train, Y_train, num_iterations, num_wolves);  

% SET THE OPTIMIZED WEIGHTS BACK INTO THE NETWORK  
if numel(alpha_pos) == totalWeights  
    final_net = setwb(final_net, alpha_pos');  % Set the best weights  
else  
    error('Size of alpha_pos does not match the expected number of weights.');  
end  

% VALIDATE THE NETWORK  
disp('Weights have been set successfully. Here are the updated weights:');  
disp(getwb(final_net)); % Display the updated weights  

% DEFINE THE GWO FUNCTION  
function alpha_pos = gwo_optimization(final_net, X_train, Y_train, num_iterations, num_wolves)  
    % Initialize the positions of wolves randomly  
    dim = numel(getwb(final_net)); % Number of weights  
    positions = rand(num_wolves, dim); % Randomly initialized positions of wolves  
    fitness_values = zeros(num_wolves, 1); % To store fitness values  

    % Ensure Y_train is a column vector for consistent comparison  
    Y_train = Y_train(:); % Convert Y_train to a column vector  

    for iter = 1:num_iterations  
        % Evaluate fitness of each wolf  
        for i = 1:num_wolves  
            % Set the weights to the network from the position  
            final_net = setwb(final_net, positions(i, :)');  

            % Calculate the output and loss (Mean Squared Error as an example)  
            Y_pred = final_net(X_train); % Transpose X_train for prediction  

            % Ensure the output dimensions match for MSE calculation  
             Y_pred = Y_pred(:); % Convert Y_pred to a column vector  
% a=Y_pred
% b=Y_train
            fitness_values(i) = mean((Y_train - Y_pred).^2); % Example fitness function: MSE  
        end  

        % Sort wolves by fitness and update alpha (best) position  
        [~, sorted_indices] = sort(fitness_values);  
        best_wolf = positions(sorted_indices(1), :);  

        % Update positions (This is a simplification of GWO)  
        a = 2 - iter * (2 / num_iterations); % Decreasing factor  
        for i = 1:num_wolves  
            for j = 1:dim  
                r = rand(); % Random number in [0,1]  
                % Update position based on best wolf (simplified)  
                positions(i, j) = best_wolf(j) + r * a * (best_wolf(j) - positions(i, j));  
            end  
        end  
    end  

    alpha_pos = best_wolf; % Return the best position found  
    
      
            
          K=12; 
      Y_train1=Y_train';
CrossValIndices = crossvalind('Kfold', 12, K);
for III = 1: K
    display(['Cross validation, folds ' num2str(III)])
    IndicesI = CrossValIndices==III;
    TempInd = CrossValIndices;
    TempInd(IndicesI) = [];
    xTraining = X_train(:,CrossValIndices~=III);
    tTrain = Y_train1(:,CrossValIndices~=III);
      xTest = X_train(:,CrossValIndices ==III);
      tTest = Y_train1(:,CrossValIndices ==III);
      
 Y_pred = final_net(xTraining);
% net = feedforwardnet([30 20 1]);
% net = train(net,xTest,tTest);

y(III,:) = final_net(xTest);
perf (III) = perform(final_net,tTest,y(III))
% perf (i) = perform(net,tTest,y(i))
tTest1(:,III)=tTest;

end     
    

act_pred_diff = [tTest1' y y-tTest1']  %display actual,predicted & error
mse=mean((tTest1' -y).^2);
disp(['Mean Squared Error: ', num2str(mse)]); 
rmse=sqrt(mse);
disp(['Root Mean Squared Error: ', num2str(rmse)]); 

plot(y,'+k','LineWidth',2),hold on, plot(tTest1,'r*')%,'LineWidth',2)
    legend('predicted','actual')

end  
